Skip to content

Commit c8704a7

Browse files
authored
Fix race in rtp.Session. (#81)
1 parent f88006a commit c8704a7

2 files changed

Lines changed: 203 additions & 25 deletions

File tree

rtp/session.go

Lines changed: 34 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -28,8 +28,7 @@ import (
2828
)
2929

3030
const (
31-
enableZeroCopy = true
32-
MTUSize = 1500
31+
MTUSize = 1500
3332
)
3433

3534
type Session interface {
@@ -168,40 +167,51 @@ type readStream struct {
168167
}
169168

170169
func (r *readStream) write(p *rtp.Packet) {
171-
if enableZeroCopy {
172-
r.mu.Lock()
173-
h, payload := r.hdr, r.payload
170+
r.mu.Lock()
171+
172+
if r.hdr != nil {
173+
// zero copy
174+
*r.hdr = p.Header
175+
n := copy(r.payload, p.Payload)
174176
r.hdr, r.payload = nil, nil
175177
r.mu.Unlock()
176-
if h != nil {
177-
// zero copy
178-
*h = p.Header
179-
n := copy(payload, p.Payload)
180-
select {
181-
case <-r.closed:
182-
case r.copied <- n:
183-
}
184-
return
178+
179+
select {
180+
case <-r.closed:
181+
case r.copied <- n:
185182
}
183+
return
186184
}
187185
p.Payload = slices.Clone(p.Payload)
188186
select {
189187
case r.recv <- p:
190188
default:
191189
}
190+
r.mu.Unlock()
192191
}
193192

194193
func (r *readStream) ReadRTP(h *rtp.Header, payload []byte) (int, error) {
195194
direct := false
196-
if enableZeroCopy {
197-
r.mu.Lock()
198-
if r.hdr == nil {
199-
r.hdr = h
200-
r.payload = payload
201-
direct = true
202-
}
195+
r.mu.Lock()
196+
197+
// Check the queue before offering our buffer.
198+
select {
199+
case p := <-r.recv:
203200
r.mu.Unlock()
201+
*h = p.Header
202+
n := copy(payload, p.Payload)
203+
return n, nil
204+
default:
204205
}
206+
207+
if r.hdr == nil {
208+
r.hdr = h
209+
r.payload = payload
210+
direct = true
211+
}
212+
r.mu.Unlock()
213+
214+
// If we didn't successfully make an offer, read from the queue.
205215
if !direct {
206216
select {
207217
case p := <-r.recv:
@@ -212,20 +222,19 @@ func (r *readStream) ReadRTP(h *rtp.Header, payload []byte) (int, error) {
212222
}
213223
return 0, io.EOF
214224
}
225+
215226
defer func() {
216227
r.mu.Lock()
217228
defer r.mu.Unlock()
218229
if r.hdr == h {
219230
r.hdr, r.payload = nil, nil
220231
}
221232
}()
233+
234+
// If we were able to offer our buffer, wait for the copy signal.
222235
select {
223236
case n := <-r.copied:
224237
return n, nil
225-
case p := <-r.recv:
226-
*h = p.Header
227-
n := copy(payload, p.Payload)
228-
return n, nil
229238
case <-r.closed:
230239
}
231240
return 0, io.EOF

rtp/session_test.go

Lines changed: 169 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,169 @@
1+
// Copyright 2026 LiveKit, Inc.
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// http://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
// See the License for the specific language governing permissions and
13+
// limitations under the License.
14+
15+
package rtp
16+
17+
import (
18+
"encoding/binary"
19+
"fmt"
20+
"net"
21+
"sync"
22+
"testing"
23+
"time"
24+
25+
"github.com/pion/rtp"
26+
"github.com/stretchr/testify/require"
27+
28+
"github.com/livekit/protocol/logger"
29+
)
30+
31+
const (
32+
testStreamSSRC = 0x11223344
33+
testSentinelSSRC = 0x55667788
34+
testPayloadType = 8 // PCMA
35+
testPayloadSize = 160 // 20ms of PCMA; a realistic size, and a wider copy() race window
36+
testClockPerPkt = 160
37+
)
38+
39+
func newRTPPacket(t *testing.T, ssrc uint32, sequenceNumber int) []byte {
40+
t.Helper()
41+
require.Less(t, sequenceNumber, 1<<16, "index must fit the 16-bit sequence number")
42+
p := rtp.Packet{
43+
Header: rtp.Header{
44+
Version: 2,
45+
PayloadType: testPayloadType,
46+
SequenceNumber: uint16(sequenceNumber),
47+
Timestamp: uint32(sequenceNumber) * testClockPerPkt,
48+
SSRC: ssrc,
49+
},
50+
Payload: payloadForSequenceNumber(sequenceNumber),
51+
}
52+
buf, err := p.Marshal()
53+
require.NoError(t, err)
54+
return buf
55+
}
56+
57+
func payloadForSequenceNumber(sequenceNumber int) []byte {
58+
payload := make([]byte, testPayloadSize)
59+
for off := 0; off < len(payload); off += 2 {
60+
binary.BigEndian.PutUint16(payload[off:], uint16(sequenceNumber))
61+
}
62+
return payload
63+
}
64+
65+
func sequenceNumberFromPayload(payload []byte) int {
66+
if len(payload) < 2 || len(payload)%2 != 0 {
67+
return -1
68+
}
69+
return int(binary.BigEndian.Uint16(payload))
70+
}
71+
72+
func verifyRTPPacket(h *rtp.Header, payload []byte) error {
73+
// Ensure that the header is properly written to.
74+
if h.Version == 0 && h.PayloadType == 0 && h.SSRC == 0 && h.SequenceNumber == 0 && h.Timestamp == 0 {
75+
return fmt.Errorf("header never written: ReadRTP reported %d bytes carrying packet %d, but left the header zeroed",
76+
len(payload), sequenceNumberFromPayload(payload))
77+
}
78+
79+
if h.Version != 2 || h.PayloadType != testPayloadType || h.SSRC != testStreamSSRC {
80+
return fmt.Errorf("torn header: got version=%d pt=%d ssrc=%#x, want version=2 pt=%d ssrc=%#x (seq=%d ts=%d)",
81+
h.Version, h.PayloadType, h.SSRC, testPayloadType, uint32(testStreamSSRC), h.SequenceNumber, h.Timestamp)
82+
}
83+
84+
sequenceNumber := int(h.SequenceNumber)
85+
if exp := uint32(sequenceNumber) * testClockPerPkt; h.Timestamp != exp {
86+
return fmt.Errorf("torn header: seq=%d implies ts=%d, got ts=%d - header fields came from two packets",
87+
sequenceNumber, exp, h.Timestamp)
88+
}
89+
90+
if len(payload) != testPayloadSize {
91+
return fmt.Errorf("unexpected payload length: header says seq=%d, but ReadRTP returned %d payload bytes, want %d",
92+
sequenceNumber, len(payload), testPayloadSize)
93+
}
94+
95+
gotSequeneceNumber := sequenceNumberFromPayload(payload)
96+
if gotSequeneceNumber != sequenceNumber {
97+
return fmt.Errorf("mixed payload: header says seq=%d, but the buffer holds packet %d",
98+
sequenceNumber, gotSequeneceNumber)
99+
}
100+
return nil
101+
}
102+
103+
func TestSessionZeroCopyHandoff(t *testing.T) {
104+
numPackets := 8000
105+
cli, srv := net.Pipe()
106+
defer cli.Close()
107+
require.NoError(t, srv.SetReadDeadline(time.Now().Add(30*time.Second)))
108+
109+
sess := NewSession(logger.GetLogger(), srv)
110+
defer sess.Close()
111+
112+
sent := make([][]byte, 0, numPackets+1)
113+
for i := range numPackets {
114+
sent = append(sent, newRTPPacket(t, testStreamSSRC, i))
115+
}
116+
sent = append(sent, newRTPPacket(t, testSentinelSSRC, 0))
117+
118+
var wg sync.WaitGroup
119+
wg.Go(func() {
120+
for _, p := range sent {
121+
if _, err := cli.Write(p); err != nil {
122+
return
123+
}
124+
}
125+
})
126+
127+
r, ssrc, err := sess.AcceptStream()
128+
require.NoError(t, err)
129+
require.EqualValues(t, uint32(testStreamSSRC), ssrc)
130+
131+
allReceived := make(chan struct{})
132+
wg.Go(func() {
133+
defer close(allReceived)
134+
for {
135+
_, ssrc, err := sess.AcceptStream()
136+
if err != nil || ssrc == testSentinelSSRC {
137+
return
138+
}
139+
}
140+
})
141+
142+
var delivered int
143+
var verifyErr error
144+
done := make(chan struct{})
145+
go func() {
146+
defer close(done)
147+
var h rtp.Header
148+
buf := make([]byte, MTUSize+1)
149+
for {
150+
h = rtp.Header{}
151+
n, err := r.ReadRTP(&h, buf)
152+
if err != nil {
153+
return
154+
}
155+
delivered++
156+
if verifyErr = verifyRTPPacket(&h, buf[:n]); verifyErr != nil {
157+
return
158+
}
159+
}
160+
}()
161+
162+
wg.Wait()
163+
sess.Close() // releases the reader with io.EOF
164+
<-done
165+
166+
require.NoError(t, verifyErr, "ReadRTP delivered a corrupted packet")
167+
require.NotZero(t, delivered, "no packets were delivered")
168+
t.Logf("delivered %d/%d packets", delivered, numPackets)
169+
}

0 commit comments

Comments
 (0)