Validate RTP packets. (#2778)

* Validate RTP packets.

Check version, payload type (if available) and SSRC (if available)
and drop bad packets. And let repair mechanisms take effect for those
packets.

* address data race reported by test

* fix an unlock and test packets
This commit is contained in:
Raja Subramanian
2024-06-10 15:43:59 +05:30
committed by GitHub
parent a31f59b689
commit 129ba62d61
5 changed files with 81 additions and 15 deletions
+1 -3
View File
@@ -358,9 +358,7 @@ func (t *MediaTrack) AddReceiver(receiver *webrtc.RTPReceiver, track *webrtc.Tra
bitrates = int(ti.Layers[layer].GetBitrate())
}
if t.IsSimulcast() {
t.MediaTrackReceiver.SetLayerSsrc(mime, track.RID(), uint32(track.SSRC()))
}
t.MediaTrackReceiver.SetLayerSsrc(mime, track.RID(), uint32(track.SSRC()))
buff.Bind(receiver.GetParameters(), track.Codec().RTPCodecCapability, bitrates)
+12 -6
View File
@@ -316,19 +316,25 @@ func (b *Buffer) Write(pkt []byte) (n int, err error) {
return
}
if rtpPacket.Version != 2 || (b.payloadType != 0 && rtpPacket.PayloadType != b.payloadType) {
if err = utils.ValidateRTPPacket(&rtpPacket, b.payloadType, b.mediaSSRC); err != nil {
b.logger.Warnw(
"invalid RTP packet", nil,
"validating RTP packet failed", err,
"version", rtpPacket.Version,
"sn", rtpPacket.SequenceNumber,
"timestamp", rtpPacket.Timestamp,
"payloadSize", len(rtpPacket.Payload),
"padding", rtpPacket.Padding,
"marker", rtpPacket.Marker,
"expectedPayloadType", b.payloadType,
"payloadType", rtpPacket.PayloadType,
"sequenceNumber", rtpPacket.SequenceNumber,
"timestamp", rtpPacket.Timestamp,
"expectedSSRC", b.mediaSSRC,
"ssrc", rtpPacket.SSRC,
"numExtensions", len(rtpPacket.Extensions),
"payloadSize", len(rtpPacket.Payload),
"rtpStats", b.rtpStats,
"snRangeMap", b.snRangeMap,
)
// TODO-REMOVE-AFTER-DEBUG
b.Unlock()
return
}
now := time.Now()
+41 -5
View File
@@ -44,7 +44,7 @@ var opusCodec = webrtc.RTPCodecParameters{
MimeType: "audio/opus",
ClockRate: 48000,
},
PayloadType: 96,
PayloadType: 111,
}
func TestNack(t *testing.T) {
@@ -81,7 +81,13 @@ func TestNack(t *testing.T) {
time.Sleep(500 * time.Millisecond) // even a long wait should not exceed max retries
}
pkt := rtp.Packet{
Header: rtp.Header{SequenceNumber: uint16(i), Timestamp: uint32(i)},
Header: rtp.Header{
Version: 2,
PayloadType: 96,
SequenceNumber: uint16(i),
Timestamp: uint32(i),
SSRC: 123,
},
Payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1},
}
b, err := pkt.Marshal()
@@ -140,7 +146,13 @@ func TestNack(t *testing.T) {
time.Sleep(500 * time.Millisecond) // even a long wait should not exceed max retries
}
pkt := rtp.Packet{
Header: rtp.Header{SequenceNumber: uint16(i + 65533), Timestamp: uint32(i)},
Header: rtp.Header{
Version: 2,
PayloadType: 96,
SequenceNumber: uint16(i + 65533),
Timestamp: uint32(i),
SSRC: 123,
},
Payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1},
}
b, err := pkt.Marshal()
@@ -166,23 +178,35 @@ func TestNewBuffer(t *testing.T) {
var TestPackets = []*rtp.Packet{
{
Header: rtp.Header{
Version: 2,
PayloadType: 96,
SequenceNumber: 65533,
SSRC: 123,
},
},
{
Header: rtp.Header{
Version: 2,
PayloadType: 96,
SequenceNumber: 65534,
SSRC: 123,
},
Payload: []byte{1},
},
{
Header: rtp.Header{
Version: 2,
PayloadType: 96,
SequenceNumber: 2,
SSRC: 123,
},
},
{
Header: rtp.Header{
Version: 2,
PayloadType: 96,
SequenceNumber: 65535,
SSRC: 123,
},
},
}
@@ -232,7 +256,13 @@ func TestFractionLostReport(t *testing.T) {
}, opusCodec.RTPCodecCapability, 0)
for i := 0; i < 15; i++ {
pkt := rtp.Packet{
Header: rtp.Header{SequenceNumber: uint16(i), Timestamp: uint32(i)},
Header: rtp.Header{
Version: 2,
PayloadType: 111,
SequenceNumber: uint16(i),
Timestamp: uint32(i),
SSRC: 123,
},
Payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1},
}
b, err := pkt.Marshal()
@@ -264,7 +294,13 @@ func TestFractionLostReport(t *testing.T) {
}, opusCodec.RTPCodecCapability, 0)
for i := 0; i < 15; i++ {
pkt := rtp.Packet{
Header: rtp.Header{SequenceNumber: uint16(i), Timestamp: uint32(i)},
Header: rtp.Header{
Version: 2,
PayloadType: 111,
SequenceNumber: uint16(i),
Timestamp: uint32(i),
SSRC: 123,
},
Payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1},
}
b, err := pkt.Marshal()
+8 -1
View File
@@ -264,7 +264,7 @@ func (s *StreamTrackerManager) RemoveAllTrackers() {
s.trackers[layer] = nil
}
s.availableLayers = make([]int32, 0)
s.maxExpectedLayerFromTrackInfo()
s.maxExpectedLayerFromTrackInfoLocked()
s.paused = false
ddTracker := s.ddTracker
s.ddTracker = nil
@@ -530,6 +530,13 @@ func (s *StreamTrackerManager) removeAvailableLayer(layer int32) {
}
func (s *StreamTrackerManager) maxExpectedLayerFromTrackInfo() {
s.lock.Lock()
defer s.lock.Unlock()
s.maxExpectedLayerFromTrackInfoLocked()
}
func (s *StreamTrackerManager) maxExpectedLayerFromTrackInfoLocked() {
s.maxExpectedLayer = buffer.InvalidLayerSpatial
ti := s.trackInfo.Load()
if ti != nil {
+19
View File
@@ -15,9 +15,11 @@
package utils
import (
"errors"
"strings"
"github.com/pion/interceptor"
"github.com/pion/rtp"
"github.com/pion/webrtc/v3"
)
@@ -51,3 +53,20 @@ func GetHeaderExtensionID(extensions []interceptor.RTPHeaderExtension, extension
}
return 0
}
// ValidateRTPPacket checks for a valid RTP packet and returns an error if fields are incorrect
func ValidateRTPPacket(pkt *rtp.Packet, expectedPayloadType uint8, expectedSSRC uint32) error {
if pkt.Version != 2 {
return errors.New("invalid RTP version")
}
if expectedPayloadType != 0 && pkt.PayloadType != expectedPayloadType {
return errors.New("invalid RTP payload type")
}
if expectedSSRC != 0 && pkt.SSRC != expectedSSRC {
return errors.New("invalid RTP SSRC")
}
return nil
}