diff --git a/pkg/rfc8888/interceptor.go b/pkg/rfc8888/interceptor.go index f513a1af..9a47e09b 100644 --- a/pkg/rfc8888/interceptor.go +++ b/pkg/rfc8888/interceptor.go @@ -6,6 +6,7 @@ package rfc8888 import ( + "errors" "sync" "time" @@ -14,6 +15,8 @@ import ( "github.com/pion/rtcp" ) +var errClosed = errors.New("interceptor is closed") + // TickerFactory is a factory to create new tickers. type TickerFactory func(d time.Duration) ticker @@ -125,7 +128,11 @@ func (s *SenderInterceptor) BindRemoteStream( sequenceNumber: header.SequenceNumber, ecn: 0, // ECN is not supported (yet). } - s.packetChan <- p + select { + case <-s.close: + return 0, nil, errClosed + case s.packetChan <- p: + } return i, attr, nil }) diff --git a/pkg/rfc8888/interceptor_test.go b/pkg/rfc8888/interceptor_test.go index 0b8b70de..f535af64 100644 --- a/pkg/rfc8888/interceptor_test.go +++ b/pkg/rfc8888/interceptor_test.go @@ -282,3 +282,44 @@ func TestInterceptor(t *testing.T) { }, ccfb.ReportBlocks[0].MetricBlocks) }) } + +func TestReadAfterClose(t *testing.T) { + f, err := NewSenderInterceptor() + assert.NoError(t, err) + + intcp, err := f.NewInterceptor("") + assert.NoError(t, err) + + intcp.BindRTCPWriter(interceptor.RTCPWriterFunc( + func(pkts []rtcp.Packet, attributes interceptor.Attributes) (int, error) { + return 0, nil + }, + )) + + raw, err := (&rtp.Packet{ + Header: rtp.Header{Version: 2, SequenceNumber: 1, SSRC: 123456}, + Payload: []byte{}, + }).Marshal() + assert.NoError(t, err) + + reader := intcp.BindRemoteStream(&interceptor.StreamInfo{SSRC: 123456}, interceptor.RTPReaderFunc( + func(b []byte, a interceptor.Attributes) (int, interceptor.Attributes, error) { + return copy(b, raw), a, nil + }, + )) + + assert.NoError(t, intcp.Close()) + + done := make(chan struct{}) + go func() { + defer close(done) + _, _, err := reader.Read(make([]byte, 1500), nil) + assert.ErrorIs(t, err, errClosed) + }() + + select { + case <-done: + case <-time.After(time.Second): + assert.Fail(t, "read after close blocked") + } +}