Skip to content

Commit 84e9935

Browse files
Nukittshubhamdhama
authored andcommitted
drpcmanager: fix data race between newStream and Close on WaitGroup
Move `wg.Add(1)` inside `activeStreams.Add` under its mutex so the increment and the closed check are atomic. This prevents `wg.Add` racing with `wg.Wait` in `Manager.Close` when terminate closes streams before `newStream` finishes registering.
1 parent b8b4dcb commit 84e9935

4 files changed

Lines changed: 55 additions & 12 deletions

File tree

drpcmanager/active_streams.go

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -25,8 +25,11 @@ func newActiveStreams() *activeStreams {
2525
}
2626

2727
// Add adds a stream. It returns an error if the collection is closed or if a
28-
// stream with the same ID already exists.
29-
func (r *activeStreams) Add(id uint64, stream *drpcstream.Stream) error {
28+
// stream with the same ID already exists. If wg is non-nil, it is incremented
29+
// under the same lock that checks closed, so wg.Add and the closed check are
30+
// atomic with respect to Close (which sets closed before Manager.Close calls
31+
// wg.Wait).
32+
func (r *activeStreams) Add(id uint64, stream *drpcstream.Stream, wg *sync.WaitGroup) error {
3033
if stream == nil {
3134
return managerClosed.New("stream can't be nil")
3235
}
@@ -41,6 +44,9 @@ func (r *activeStreams) Add(id uint64, stream *drpcstream.Stream) error {
4144
return managerClosed.New("duplicate stream id")
4245
}
4346
r.streams[id] = stream
47+
if wg != nil {
48+
wg.Add(1)
49+
}
4450
return nil
4551
}
4652

drpcmanager/active_streams_test.go

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ func TestActiveStreams_AddAndGet(t *testing.T) {
2929
streams := newActiveStreams()
3030
s := testStream(t, 1)
3131

32-
assert.NoError(t, streams.Add(1, s))
32+
assert.NoError(t, streams.Add(1, s, nil))
3333

3434
got, ok := streams.Get(1)
3535
assert.That(t, ok)
@@ -48,7 +48,7 @@ func TestActiveStreams_Remove(t *testing.T) {
4848
streams := newActiveStreams()
4949
s := testStream(t, 1)
5050

51-
assert.NoError(t, streams.Add(1, s))
51+
assert.NoError(t, streams.Add(1, s, nil))
5252
assert.Equal(t, streams.Len(), 1)
5353

5454
streams.Remove(1)
@@ -70,8 +70,8 @@ func TestActiveStreams_DuplicateAdd(t *testing.T) {
7070
s1 := testStream(t, 1)
7171
s2 := testStream(t, 1)
7272

73-
assert.NoError(t, streams.Add(1, s1))
74-
assert.Error(t, streams.Add(1, s2))
73+
assert.NoError(t, streams.Add(1, s1, nil))
74+
assert.Error(t, streams.Add(1, s2, nil))
7575

7676
// original stream is still present
7777
got, ok := streams.Get(1)
@@ -83,14 +83,14 @@ func TestActiveStreams_AddAfterClose(t *testing.T) {
8383
streams := newActiveStreams()
8484
streams.Close(errors.New("closed"))
8585

86-
err := streams.Add(1, testStream(t, 1))
86+
err := streams.Add(1, testStream(t, 1), nil)
8787
assert.Error(t, err)
8888
}
8989

9090
func TestActiveStreams_RemoveAfterClose(t *testing.T) {
9191
streams := newActiveStreams()
9292
s := testStream(t, 1)
93-
assert.NoError(t, streams.Add(1, s))
93+
assert.NoError(t, streams.Add(1, s, nil))
9494

9595
streams.Close(errors.New("closed"))
9696

@@ -102,10 +102,10 @@ func TestActiveStreams_Len(t *testing.T) {
102102
streams := newActiveStreams()
103103
assert.Equal(t, streams.Len(), 0)
104104

105-
assert.NoError(t, streams.Add(1, testStream(t, 1)))
105+
assert.NoError(t, streams.Add(1, testStream(t, 1), nil))
106106
assert.Equal(t, streams.Len(), 1)
107107

108-
assert.NoError(t, streams.Add(2, testStream(t, 2)))
108+
assert.NoError(t, streams.Add(2, testStream(t, 2), nil))
109109
assert.Equal(t, streams.Len(), 2)
110110

111111
streams.Remove(1)

drpcmanager/manager.go

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -324,15 +324,14 @@ func (m *Manager) newStream(ctx context.Context, sid uint64, kind drpc.StreamKin
324324

325325
stream := drpcstream.NewWithOptions(ctx, sid, m.wr, m.recvPool, opts)
326326

327-
if err := m.streams.Add(sid, stream); err != nil {
327+
if err := m.streams.Add(sid, stream, &m.wg); err != nil {
328328
return nil, err
329329
}
330330

331331
if m.metrics.ShouldRecord() {
332332
m.metrics.StreamsStarted.Inc(1)
333333
}
334334

335-
m.wg.Add(1)
336335
go m.manageStream(ctx, stream)
337336

338337
m.log("STREAM", stream.String)

drpcmanager/manager_test.go

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -588,3 +588,41 @@ func TestManager_ServerClientHangupCancels(t *testing.T) {
588588
assert.That(t, errors.Is(err, context.Canceled))
589589
assert.Equal(t, status.Code(drpc.ToRPCErr(err)), codes.Canceled)
590590
}
591+
592+
// TestManager_ConcurrentCloseAndNewClientStream exercises the race between
593+
// Manager.Close (terminate → stop writer, close transport, close streams) and
594+
// Manager.NewClientStream (newStream → create stream, add to streams, wg.Add).
595+
// Under -race this would fail without wg.Add being atomic with the closed
596+
// check in activeStreams.Add.
597+
func TestManager_ConcurrentCloseAndNewClientStream(t *testing.T) {
598+
for i := 0; i < 100; i++ {
599+
cconn, sconn := net.Pipe()
600+
601+
cman := New(cconn, Client)
602+
603+
start := make(chan struct{})
604+
done := make(chan struct{}, 2)
605+
606+
go func() {
607+
<-start
608+
_ = cman.Close()
609+
done <- struct{}{}
610+
}()
611+
612+
go func() {
613+
<-start
614+
stream, err := cman.NewClientStream(context.Background(), "rpc", 0)
615+
if err == nil {
616+
_ = stream.Close()
617+
}
618+
done <- struct{}{}
619+
}()
620+
621+
close(start)
622+
<-done
623+
<-done
624+
625+
_ = cconn.Close()
626+
_ = sconn.Close()
627+
}
628+
}

0 commit comments

Comments
 (0)