From a8b35cc8186fc79f4e54550d30f8f0b565247bbd Mon Sep 17 00:00:00 2001 From: David Chen Date: Sat, 29 Aug 2026 22:02:23 -0700 Subject: [PATCH] Harden upstream FlexFEC handling --- pkg/rtc/config.go | 12 +- pkg/rtc/transport.go | 12 +- pkg/rtc/transport_fec_test.go | 145 ++++++++++++--- pkg/sfu/buffer/buffer.go | 29 ++- pkg/sfu/buffer/buffer_fec_test.go | 69 +++++++- pkg/sfu/flexfec/decoder.go | 72 +++----- pkg/sfu/flexfec/decoder_benchmark_test.go | 10 +- pkg/sfu/flexfec/decoder_test.go | 204 ++++++++++++++++++---- test/flexfec_upstream_test.go | 3 + 9 files changed, 421 insertions(+), 135 deletions(-) diff --git a/pkg/rtc/config.go b/pkg/rtc/config.go index 37ea20330..5d856b1c5 100644 --- a/pkg/rtc/config.go +++ b/pkg/rtc/config.go @@ -70,6 +70,12 @@ type FlexFECDirectionConfig struct { func NewWebRTCConfig(conf *config.Config) (*WebRTCConfig, error) { rtcConf := conf.RTC + flexFEC := rtcConf.FlexFEC.WithDefaults() + if flexFEC.UpstreamEnabled { + if err := validateFlexFECPayloadType(flexFEC.PayloadType); err != nil { + return nil, err + } + } webRTCConfig, err := rtcconfig.NewWebRTCConfig(&rtcConf.RTCConfig, conf.Development) if err != nil { @@ -89,12 +95,6 @@ func NewWebRTCConfig(conf *config.Config) (*WebRTCConfig, error) { rtcConf.PacketBufferSizeAudio = rtcConf.PacketBufferSize } - flexFEC := rtcConf.FlexFEC.WithDefaults() - if flexFEC.UpstreamEnabled { - if err := validateFlexFECPayloadType(flexFEC.PayloadType); err != nil { - return nil, err - } - } return &WebRTCConfig{ WebRTCConfig: *webRTCConfig, Receiver: ReceiverConfig{ diff --git a/pkg/rtc/transport.go b/pkg/rtc/transport.go index 7a9e4ffe5..5bf75c9b8 100644 --- a/pkg/rtc/transport.go +++ b/pkg/rtc/transport.go @@ -3407,18 +3407,18 @@ func fecPairsFromSDP(s *sdp.SessionDescription, logger logger.Logger) map[uint32 if attr.Key != sdp.AttrKeySSRCGroup { continue } - split := strings.Split(attr.Value, " ") - if split[0] != sdp.SemanticTokenForwardErrorCorrectionFramework || len(split) != 3 { + fields := strings.Fields(attr.Value) + if len(fields) != 3 || fields[0] != sdp.SemanticTokenForwardErrorCorrectionFramework { continue } - baseSsrc, err := strconv.ParseUint(split[1], 10, 32) + baseSsrc, err := strconv.ParseUint(fields[1], 10, 32) if err != nil { - logger.Warnw("failed to parse SSRC", err, "ssrc", split[1]) + logger.Warnw("failed to parse SSRC", err, "ssrc", fields[1]) continue } - fecSsrc, err := strconv.ParseUint(split[2], 10, 32) + fecSsrc, err := strconv.ParseUint(fields[2], 10, 32) if err != nil { - logger.Warnw("failed to parse SSRC", err, "ssrc", split[2]) + logger.Warnw("failed to parse SSRC", err, "ssrc", fields[2]) continue } fecFlows[uint32(fecSsrc)] = uint32(baseSsrc) diff --git a/pkg/rtc/transport_fec_test.go b/pkg/rtc/transport_fec_test.go index cd7763715..fb33b66f0 100644 --- a/pkg/rtc/transport_fec_test.go +++ b/pkg/rtc/transport_fec_test.go @@ -22,6 +22,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/livekit/livekit-server/pkg/config" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" ) @@ -76,6 +77,37 @@ a=ssrc:1111 cname:test assert.Empty(t, fecPairsFromSDP(parsed, logger.GetLogger())) } +func TestFECPairsFromSDPIgnoresMalformedGroups(t *testing.T) { + description := &sdp.SessionDescription{ + MediaDescriptions: []*sdp.MediaDescription{{ + Attributes: []sdp.Attribute{ + {Key: sdp.AttrKeySSRCGroup}, + {Key: sdp.AttrKeySSRCGroup, Value: "FEC-FR 1111"}, + {Key: sdp.AttrKeySSRCGroup, Value: "FEC-FR invalid 3333"}, + {Key: sdp.AttrKeySSRCGroup, Value: "FEC-FR 1111 invalid"}, + {Key: sdp.AttrKeySSRCGroup, Value: "FEC-FR 1111 3333 4444"}, + }, + }}, + } + + require.NotPanics(t, func() { + assert.Empty(t, fecPairsFromSDP(description, logger.GetLogger())) + }) +} + +func TestFECPairsFromSDPHandlesWhitespace(t *testing.T) { + description := &sdp.SessionDescription{ + MediaDescriptions: []*sdp.MediaDescription{{ + Attributes: []sdp.Attribute{{ + Key: sdp.AttrKeySSRCGroup, + Value: " FEC-FR 1111\t3333 ", + }}, + }}, + } + + assert.Equal(t, map[uint32]uint32{3333: 1111}, fecPairsFromSDP(description, logger.GetLogger())) +} + func TestFlexFECPayloadTypeValidation(t *testing.T) { assert.NoError(t, validateFlexFECPayloadType(115)) // upper boundary of the 7-bit RTP payload type field @@ -96,35 +128,90 @@ func TestMediaEngineRegistersFlexFEC(t *testing.T) { {Mime: "video/rtx"}, } - for _, enabled := range []bool{false, true} { - me, err := createMediaEngine(enabledCodecs, DirectionConfig{ - FlexFEC: FlexFECDirectionConfig{ - Enabled: enabled, - PayloadType: 115, - }, - RTCPFeedback: RTCPFeedbackConfig{ - Video: []webrtc.RTCPFeedback{{Type: webrtc.TypeRTCPFBTransportCC}}, - }, - }, false) - require.NoError(t, err) + for _, test := range []struct { + name string + enabled bool + }{ + {name: "disabled", enabled: false}, + {name: "enabled", enabled: true}, + } { + t.Run(test.name, func(t *testing.T) { + me, err := createMediaEngine(enabledCodecs, DirectionConfig{ + FlexFEC: FlexFECDirectionConfig{ + Enabled: test.enabled, + PayloadType: 115, + }, + RTCPFeedback: RTCPFeedbackConfig{ + Video: []webrtc.RTCPFeedback{{Type: webrtc.TypeRTCPFBTransportCC}}, + }, + }, false) + require.NoError(t, err) - // drive codec registration into negotiated form via an SDP round trip - // is heavyweight, instead check via filterCodecs retention behavior - flexfecParams := flexFECCodecParameters(115) - assert.Equal(t, "repair-window=2000000", flexfecParams.SDPFmtpLine) - filtered := filterCodecs( - []webrtc.RTPCodecParameters{flexfecParams}, - enabledCodecs, - RTCPFeedbackConfig{}, - false, - enabled, - ) - if enabled { - require.Len(t, filtered, 1) - assert.Equal(t, webrtc.MimeTypeFlexFEC03, filtered[0].MimeType) - } else { - assert.Empty(t, filtered) - } - _ = me + pc, err := webrtc.NewAPI(webrtc.WithMediaEngine(me)).NewPeerConnection(webrtc.Configuration{}) + require.NoError(t, err) + defer pc.Close() + _, err = pc.AddTransceiverFromKind(webrtc.RTPCodecTypeVideo) + require.NoError(t, err) + offer, err := pc.CreateOffer(nil) + require.NoError(t, err) + + flexFECParams := flexFECCodecParameters(115) + assert.Equal(t, "repair-window=2000000", flexFECParams.SDPFmtpLine) + filtered := filterCodecs( + []webrtc.RTPCodecParameters{flexFECParams}, + enabledCodecs, + RTCPFeedbackConfig{}, + false, + test.enabled, + ) + if test.enabled { + require.Len(t, filtered, 1) + assert.Equal(t, webrtc.MimeTypeFlexFEC03, filtered[0].MimeType) + assert.Contains(t, offer.SDP, "a=rtpmap:115 flexfec-03/90000") + assert.Contains(t, offer.SDP, "a=fmtp:115 repair-window=2000000") + } else { + assert.Empty(t, filtered) + assert.NotContains(t, offer.SDP, "flexfec-03") + } + }) } } + +func TestWebRTCConfigFlexFEC(t *testing.T) { + newConfig := func(t *testing.T) *config.Config { + t.Helper() + conf, err := config.NewConfig("", true, nil, nil) + require.NoError(t, err) + conf.RTC.TCPPort = 0 + return conf + } + + t.Run("defaults and publisher updates", func(t *testing.T) { + conf := newConfig(t) + conf.RTC.FlexFEC = config.FlexFECConfig{UpstreamEnabled: true} + + webRTCConfig, err := NewWebRTCConfig(conf) + require.NoError(t, err) + assert.Equal(t, FlexFECDirectionConfig{ + Enabled: true, + PayloadType: config.DefaultFlexFECConfig.PayloadType, + }, webRTCConfig.Publisher.FlexFEC) + assert.False(t, webRTCConfig.Subscriber.FlexFEC.Enabled) + + webRTCConfig.UpdatePublisherConfig(true) + assert.True(t, webRTCConfig.Publisher.FlexFEC.Enabled) + assert.Equal(t, config.DefaultFlexFECConfig.PayloadType, webRTCConfig.Publisher.FlexFEC.PayloadType) + }) + + t.Run("invalid payload type", func(t *testing.T) { + conf := newConfig(t) + conf.RTC.FlexFEC = config.FlexFECConfig{ + UpstreamEnabled: true, + PayloadType: 96, + } + + _, err := NewWebRTCConfig(conf) + require.Error(t, err) + assert.Contains(t, err.Error(), "collides") + }) +} diff --git a/pkg/sfu/buffer/buffer.go b/pkg/sfu/buffer/buffer.go index 950a782a1..b77f24642 100644 --- a/pkg/sfu/buffer/buffer.go +++ b/pkg/sfu/buffer/buffer.go @@ -430,8 +430,8 @@ func (b *Buffer) SetPrimaryBufferForFEC(primaryBuffer *Buffer) { ssrc := b.BufferBase.SSRC() b.Unlock() - // let the primary know the repair stream SSRC so its decoder starts - // filling with media before the first FEC packet shows up + // Let the primary know the repair stream SSRC so its decoder is ready + // before the first FEC packet arrives. primaryBuffer.setFECSSRC(ssrc) for _, pp := range pkts { @@ -486,17 +486,27 @@ func (b *Buffer) maybeCreateFECDecoderLocked() { } func (b *Buffer) getFECMediaPacketLocked(sequenceNumber uint16, dst []byte) (int, error) { - if b.bucket == nil { + if b.bucket == nil || b.rtpStats == nil { return 0, errFECMediaPacketNotFound } - headSequenceNumber := b.bucket.HeadSequenceNumber() - extendedSequenceNumber := int64(headSequenceNumber) + int64(int16(sequenceNumber-uint16(headSequenceNumber))) + // FlexFEC masks use the publisher's sequence-number space. BufferBase + // removes padding-only packets from the downstream space, so resolve the + // original extended sequence number and apply the same adjustment used + // when the packet was inserted into the bucket. + highestSequenceNumber := b.rtpStats.ExtendedHighestSequenceNumber() + extendedSequenceNumber := int64(highestSequenceNumber) + int64(int16(sequenceNumber-uint16(highestSequenceNumber))) if extendedSequenceNumber < 0 { return 0, errFECMediaPacketNotFound } - return b.bucket.GetPacket(dst, uint64(extendedSequenceNumber)) + extendedSN := uint64(extendedSequenceNumber) + sequenceNumberAdjustment, err := b.snRangeMap.GetValue(extendedSN) + if err != nil || sequenceNumberAdjustment > extendedSN { + return 0, errFECMediaPacketNotFound + } + + return b.bucket.GetPacket(dst, extendedSN-sequenceNumberAdjustment) } // OnFECRecovery is called with counter deltas whenever FEC packets are @@ -567,9 +577,9 @@ func (b *Buffer) feedFECLocked( arrivalTime int64, ) (fecRecoveryDelta, func(received int, recovered int, discarded int, bytesReceived int)) { statsBefore := b.fecDecoder.Stats() - recovered := b.fecDecoder.DecodeFec(pkt) + recovered := b.fecDecoder.DecodeFEC(pkt) - if b.fecPktBuf == nil { + if len(recovered) > 0 && b.fecPktBuf == nil { b.fecPktBuf = make([]byte, bucket.RTPMaxPktSize) } for _, rp := range recovered { @@ -581,7 +591,8 @@ func (b *Buffer) feedFECLocked( // recovered packets flow through the regular pipeline: they are // forwarded downstream and stop NACKs for the lost sequence numbers. - // they do not re-enter the decoder, it already has them in its window. + // They do not re-enter the decoder because chained recovery already + // completed within DecodeFEC. b.calc(b.fecPktBuf[:n], rp, arrivalTime, false, true) } diff --git a/pkg/sfu/buffer/buffer_fec_test.go b/pkg/sfu/buffer/buffer_fec_test.go index 71766e146..df420d489 100644 --- a/pkg/sfu/buffer/buffer_fec_test.go +++ b/pkg/sfu/buffer/buffer_fec_test.go @@ -15,6 +15,7 @@ package buffer import ( + "encoding/binary" "math/rand" "testing" "time" @@ -35,7 +36,7 @@ const ( fecTestFECPT = uint8(115) ) -var flexfecCodec = webrtc.RTPCodecParameters{ +var flexFECCodec = webrtc.RTPCodecParameters{ RTPCodecCapability: webrtc.RTPCodecCapability{ MimeType: webrtc.MimeTypeFlexFEC03, ClockRate: 90000, @@ -76,7 +77,7 @@ func bindFECTestBuffer(t *testing.T, buff *Buffer) { t.Helper() buff.codecType = webrtc.RTPCodecTypeVideo require.NoError(t, buff.Bind(webrtc.RTPParameters{ - Codecs: []webrtc.RTPCodecParameters{vp8Codec, flexfecCodec}, + Codecs: []webrtc.RTPCodecParameters{vp8Codec, flexFECCodec}, }, vp8Codec.RTPCodecCapability, 0)) } @@ -386,6 +387,70 @@ func TestBufferFECSequenceNumberWrap(t *testing.T) { requireRecoveredInBucket(t, primary, &media[droppedIdx], extSNBySN, media[0].SequenceNumber) } +func TestBufferFECRecoveryAfterPaddingRemoval(t *testing.T) { + factory := NewFactoryOfBufferFactory(500, 200).CreateBufferFactory() + factory.SetFECPair(fecTestFECSSRC, fecTestMediaSSRC) + + primary := factory.GetOrNew(packetio.RTPBufferPacket, fecTestMediaSSRC).(*Buffer) + fecBuff := factory.GetOrNew(packetio.RTPBufferPacket, fecTestFECSSRC).(*Buffer) + bindFECTestBuffer(t, primary) + + media := fecTestMediaPackets(t, 800, 5) + fecPackets := pionflexfec.NewFlexEncoder03(fecTestFECPT, fecTestFECSSRC).EncodeFec(media, 1) + require.Len(t, fecPackets, 1) + + // Insert a padding-only packet into the publisher sequence-number space. + // The FEC packet protects the five media packets but not the padding packet. + for i := 1; i < len(media); i++ { + media[i].SequenceNumber++ + } + mask := uint16(0x8000) + for _, offset := range []uint{0, 2, 3, 4, 5} { + mask |= 1 << (14 - offset) + } + binary.BigEndian.PutUint16(fecPackets[0].Payload[18:20], mask) + + padding := rtp.Packet{ + Header: rtp.Header{ + Version: 2, + Padding: true, + PaddingSize: 20, + PayloadType: uint8(vp8Codec.PayloadType), + SequenceNumber: 801, + Timestamp: media[0].Timestamp, + SSRC: fecTestMediaSSRC, + }, + } + + writePacket(t, primary, &media[0]) + writePacket(t, primary, &padding) + const droppedIdx = 2 + for i := 1; i < len(media); i++ { + if i != droppedIdx { + writePacket(t, primary, &media[i]) + } + } + writePacket(t, fecBuff, &fecPackets[0]) + + require.EqualValues(t, 1, primary.FECDecoderStats().PacketsRecovered) + extSNBySN := readExtSequenceNumbers(t, primary, len(media)-1) + baseExtSN, ok := extSNBySN[media[0].SequenceNumber] + require.True(t, ok) + + // The removed padding packet shifts the recovered packet down by one in + // the bucket/downstream sequence-number space. + var raw [1500]byte + n, err := primary.GetPacket(raw[:], baseExtSN+2) + require.NoError(t, err) + var recovered rtp.Packet + require.NoError(t, recovered.Unmarshal(raw[:n])) + assert.Equal(t, media[droppedIdx].SequenceNumber-1, recovered.SequenceNumber) + assert.Equal(t, media[droppedIdx].Timestamp, recovered.Timestamp) + assert.Equal(t, media[droppedIdx].PayloadType, recovered.PayloadType) + assert.Equal(t, media[droppedIdx].SSRC, recovered.SSRC) + assert.Equal(t, media[droppedIdx].Payload, recovered.Payload) +} + func TestBufferFECIgnoresUnexpectedPayloadType(t *testing.T) { factory := NewFactoryOfBufferFactory(500, 200).CreateBufferFactory() diff --git a/pkg/sfu/flexfec/decoder.go b/pkg/sfu/flexfec/decoder.go index b5fe2191b..ae9bbb8ae 100644 --- a/pkg/sfu/flexfec/decoder.go +++ b/pkg/sfu/flexfec/decoder.go @@ -132,10 +132,10 @@ func (d *Decoder) Stats() DecoderStats { return d.stats } -// DecodeFec ingests a packet of either the FEC stream (fecSSRC) or the +// DecodeFEC ingests a packet of either the FEC stream (fecSSRC) or the // protected media stream (protectedSSRC) and returns any media packets that // became recoverable. Ownership of returned packets transfers to the caller. -func (d *Decoder) DecodeFec(receivedPacket *rtp.Packet) []*rtp.Packet { +func (d *Decoder) DecodeFEC(receivedPacket *rtp.Packet) []*rtp.Packet { switch receivedPacket.SSRC { case d.fecSSRC: d.stats.FECPacketsReceived++ @@ -172,15 +172,13 @@ func (d *Decoder) observeMediaPacket(sequenceNumber uint16) { } func (d *Decoder) discardOldFECPackets(sequenceNumber uint16) { - // Discard old FEC packets such that the sequence numbers in - // `receivedFECPackets` span at most 1/2 of the sequence number space. - // This is important for keeping `receivedFECPackets` sorted, and may - // also reduce the possibility of incorrect decoding due to sequence - // number wrap-around. + // Keep the retained sequence-number span well below half of the sequence + // space. This keeps ordering unambiguous across wrap-around and reduces the + // possibility of decoding against stale state. if len(d.receivedFECPackets) > 0 { toRemove := 0 for _, fecPkt := range d.receivedFECPackets { - if absInt(int(sequenceNumber)-int(fecPkt.packet.SequenceNumber)) > 0x3fff { + if seqDiff(sequenceNumber, fecPkt.packet.SequenceNumber) > 0x3fff { toRemove++ } else { // no need to keep iterating, since receivedFECPackets is sorted @@ -220,10 +218,7 @@ func (d *Decoder) insertFECPacket(fecPkt *rtp.Packet) { return } - var protectedSeqBuf [maxProtectedPackets]uint16 - protectedSeqs := fec.protectedSequences(protectedSeqBuf[:0]) - - if len(protectedSeqs) == 0 { + if fec.mask0 == 0 && fec.mask1 == 0 && fec.mask2 == 0 { d.stats.FECPacketsDiscarded++ d.logger.Debugw("flexfec: discarding packet", "error", errEmptyMask) return @@ -233,18 +228,13 @@ func (d *Decoder) insertFECPacket(fecPkt *rtp.Packet) { return } - // The caller may reuse packet memory after DecodeFec returns. Take + // The caller may reuse packet memory after DecodeFEC returns. Take // ownership only now that this FEC state needs to be retained. ownedFECPkt := fecPkt.Clone() - ownedFEC, err := parseFlexFEC03Header(ownedFECPkt.Payload) - if err != nil { - // Parsing the same bytes succeeded above, so this should be unreachable. - d.stats.FECPacketsDiscarded++ - d.logger.Debugw("flexfec: failed to parse cloned header", "error", err) - return - } + ownedFEC := fec + ownedFEC.payload = ownedFECPkt.Payload[len(fecPkt.Payload)-len(fec.payload):] - state := fecPacketState{packet: ownedFECPkt, flexFec: ownedFEC} + state := fecPacketState{packet: ownedFECPkt, flexFEC: ownedFEC} d.receivedFECPackets = append(d.receivedFECPackets, state) if len(d.receivedFECPackets) > 1 && !isNewerSeq( d.receivedFECPackets[len(d.receivedFECPackets)-2].packet.SequenceNumber, @@ -268,7 +258,7 @@ func (d *Decoder) attemptRecovery() []*rtp.Packet { packetsRecovered := 0 for i := 0; i < len(d.receivedFECPackets); { fecPkt := &d.receivedFECPackets[i] - packetsMissing := d.countMissingPackets(fecPkt.flexFec, recoveredPackets) + packetsMissing := d.countMissingPackets(fecPkt.flexFEC, recoveredPackets) if packetsMissing == 0 { d.removeFECPacketAt(i) continue @@ -299,7 +289,7 @@ func (d *Decoder) attemptRecovery() []*rtp.Packet { return recoveredPackets } -func (d *Decoder) countMissingPackets(fec flexFec, recoveredPackets []*rtp.Packet) int { +func (d *Decoder) countMissingPackets(fec flexFEC, recoveredPackets []*rtp.Packet) int { var protectedSeqBuf [maxProtectedPackets]uint16 protectedSeqs := fec.protectedSequences(protectedSeqBuf[:0]) missing := 0 @@ -347,7 +337,7 @@ func (d *Decoder) recoverPacket(fec *fecPacketState, recoveredPackets []*rtp.Pac var headerRecovery [12]byte copy(headerRecovery[:], fec.packet.Payload[:10]) var protectedSeqBuf [maxProtectedPackets]uint16 - protectedSeqs := fec.flexFec.protectedSequences(protectedSeqBuf[:0]) + protectedSeqs := fec.flexFEC.protectedSequences(protectedSeqBuf[:0]) missing := 0 var sequenceNumber uint16 @@ -391,7 +381,7 @@ func (d *Decoder) recoverPacket(fec *fecPacketState, recoveredPackets []*rtp.Pac recoveredRaw := make([]byte, 12+int(payloadLength)) copy(recoveredRaw[:12], headerRecovery[:]) - copy(recoveredRaw[12:], fec.flexFec.payload) + copy(recoveredRaw[12:], fec.flexFEC.payload) for _, protectedSeq := range protectedSeqs { n, err := d.getMediaPacket(protectedSeq, recoveredPackets, d.mediaPacketBuf[:]) if err != nil { @@ -425,10 +415,10 @@ func appendMaskSequences(dst []uint16, mask uint64, bitCount uint16, seqNumBase type fecPacketState struct { packet *rtp.Packet - flexFec flexFec + flexFEC flexFEC } -type flexFec struct { +type flexFEC struct { protectedSSRC uint32 seqNumBase uint16 mask0 uint16 @@ -437,7 +427,7 @@ type flexFec struct { payload []byte } -func (f flexFec) protectedSequences(dst []uint16) []uint16 { +func (f flexFEC) protectedSequences(dst []uint16) []uint16 { dst = appendMaskSequences(dst, uint64(f.mask0), fecMask0Bits, f.seqNumBase) if f.mask1 != 0 { dst = appendMaskSequences(dst, uint64(f.mask1), fecMask1Bits, f.seqNumBase+fecMask0Bits) @@ -449,24 +439,24 @@ func (f flexFec) protectedSequences(dst []uint16) []uint16 { return dst } -func parseFlexFEC03Header(data []byte) (flexFec, error) { +func parseFlexFEC03Header(data []byte) (flexFEC, error) { if len(data) < 20 { - return flexFec{}, fmt.Errorf("%w: length %d", errPacketTruncated, len(data)) + return flexFEC{}, fmt.Errorf("%w: length %d", errPacketTruncated, len(data)) } rBit := (data[0] & fecRetransmissionBit) != 0 if rBit { - return flexFec{}, errRetransmissionBitSet + return flexFEC{}, errRetransmissionBitSet } fBit := (data[0] & fecInflexibleBit) != 0 if fBit { - return flexFec{}, errInflexibleGeneratorMatrix + return flexFEC{}, errInflexibleGeneratorMatrix } ssrcCount := data[8] if ssrcCount != 1 { - return flexFec{}, fmt.Errorf("%w: count %d", errMultipleSSRCProtection, ssrcCount) + return flexFEC{}, fmt.Errorf("%w: count %d", errMultipleSSRCProtection, ssrcCount) } protectedSSRC := binary.BigEndian.Uint32(data[12:]) @@ -483,7 +473,7 @@ func parseFlexFEC03Header(data []byte) (flexFec, error) { payload = rawPacketMask[2:] } else { if len(data) < 24 { - return flexFec{}, fmt.Errorf("%w: length %d", errPacketTruncated, len(data)) + return flexFEC{}, fmt.Errorf("%w: length %d", errPacketTruncated, len(data)) } kBit1 := (rawPacketMask[2] & fecMaskKBit) != 0 @@ -493,7 +483,7 @@ func parseFlexFEC03Header(data []byte) (flexFec, error) { payload = rawPacketMask[6:] } else { if len(data) < 32 { - return flexFec{}, fmt.Errorf("%w: length %d", errPacketTruncated, len(data)) + return flexFEC{}, fmt.Errorf("%w: length %d", errPacketTruncated, len(data)) } kBit2 := (rawPacketMask[6] & fecMaskKBit) != 0 @@ -502,12 +492,12 @@ func parseFlexFEC03Header(data []byte) (flexFec, error) { if kBit2 { payload = rawPacketMask[14:] } else { - return flexFec{}, errLastOptionalMaskKBitSetToFalse + return flexFEC{}, errLastOptionalMaskKBitSetToFalse } } } - return flexFec{ + return flexFEC{ protectedSSRC: protectedSSRC, seqNumBase: seqNumBase, mask0: maskPart0, @@ -521,14 +511,6 @@ func seqDiff(a, b uint16) uint16 { return min(a-b, b-a) } -func absInt(x int) int { - if x >= 0 { - return x - } - - return -x -} - func isNewerSeq(prevValue, value uint16) bool { // half-way mark breakpoint := uint16(0x8000) diff --git a/pkg/sfu/flexfec/decoder_benchmark_test.go b/pkg/sfu/flexfec/decoder_benchmark_test.go index 99a9656fd..15fb78819 100644 --- a/pkg/sfu/flexfec/decoder_benchmark_test.go +++ b/pkg/sfu/flexfec/decoder_benchmark_test.go @@ -71,7 +71,7 @@ func BenchmarkDecoderMediaSteadyState1200(b *testing.B) { for i := 0; i < b.N; i++ { packet.SequenceNumber = uint16(i) packet.Timestamp = uint32(i) * 3000 - decoder.DecodeFec(&packet) + decoder.DecodeFEC(&packet) } } @@ -85,9 +85,9 @@ func BenchmarkDecoderCompleteWindow10x1200(b *testing.B) { for i := 0; i < b.N; i++ { decoder := NewDecoder(testFECSSRC, testMediaSSRC, lookup, logger.GetLogger()) for j := range media { - decoder.DecodeFec(&media[j]) + decoder.DecodeFEC(&media[j]) } - if recovered := decoder.DecodeFec(&fecPackets[0]); len(recovered) != 0 { + if recovered := decoder.DecodeFEC(&fecPackets[0]); len(recovered) != 0 { b.Fatalf("expected no recovered packets, got %d", len(recovered)) } } @@ -104,10 +104,10 @@ func BenchmarkDecoderRecoveryWindow10x1200(b *testing.B) { decoder := NewDecoder(testFECSSRC, testMediaSSRC, lookup, logger.GetLogger()) for j := range media { if j != 4 { - decoder.DecodeFec(&media[j]) + decoder.DecodeFEC(&media[j]) } } - if recovered := decoder.DecodeFec(&fecPackets[0]); len(recovered) != 1 { + if recovered := decoder.DecodeFEC(&fecPackets[0]); len(recovered) != 1 { b.Fatalf("expected one recovered packet, got %d", len(recovered)) } } diff --git a/pkg/sfu/flexfec/decoder_test.go b/pkg/sfu/flexfec/decoder_test.go index 22d227b9d..2946df6a7 100644 --- a/pkg/sfu/flexfec/decoder_test.go +++ b/pkg/sfu/flexfec/decoder_test.go @@ -60,7 +60,7 @@ func (d *testDecoder) getMediaPacket(sequenceNumber uint16, dst []byte) (int, er return copy(dst, packet), nil } -func (d *testDecoder) DecodeFec(packet *rtp.Packet) []*rtp.Packet { +func (d *testDecoder) DecodeFEC(packet *rtp.Packet) []*rtp.Packet { if packet.SSRC == d.protectedSSRC { raw, err := packet.Marshal() if err != nil { @@ -69,7 +69,7 @@ func (d *testDecoder) DecodeFec(packet *rtp.Packet) []*rtp.Packet { d.mediaPackets[packet.SequenceNumber] = raw } - recovered := d.Decoder.DecodeFec(packet) + recovered := d.Decoder.DecodeFEC(packet) for _, recoveredPacket := range recovered { raw, err := recoveredPacket.Marshal() if err != nil { @@ -132,12 +132,12 @@ func TestDecoderRecoversSingleLoss(t *testing.T) { if i == 2 { continue } - recovered = append(recovered, decoder.DecodeFec(&media[i])...) + recovered = append(recovered, decoder.DecodeFEC(&media[i])...) } require.Empty(t, recovered) for i := range fec { - recovered = append(recovered, decoder.DecodeFec(&fec[i])...) + recovered = append(recovered, decoder.DecodeFEC(&fec[i])...) } require.Len(t, recovered, 1) @@ -161,10 +161,10 @@ func TestDecoderRecoversPacketWithExtendedHeader(t *testing.T) { for i := range media { if i != 2 { - require.Empty(t, decoder.DecodeFec(&media[i])) + require.Empty(t, decoder.DecodeFEC(&media[i])) } } - recovered := decoder.DecodeFec(&fec[0]) + recovered := decoder.DecodeFEC(&fec[0]) require.Len(t, recovered, 1) expectedRaw, err := media[2].Marshal() @@ -184,17 +184,17 @@ func TestDecoderRecoversWithLateMedia(t *testing.T) { var recovered []*rtp.Packet for _, i := range []int{0, 3, 4} { - recovered = append(recovered, decoder.DecodeFec(&media[i])...) + recovered = append(recovered, decoder.DecodeFEC(&media[i])...) } for i := range fec { - recovered = append(recovered, decoder.DecodeFec(&fec[i])...) + recovered = append(recovered, decoder.DecodeFEC(&fec[i])...) } // two packets missing from the protected window, nothing recoverable yet require.Empty(t, recovered) require.Len(t, decoder.receivedFECPackets, 1) // late arrival of media[1] leaves only media[2] missing - recovered = decoder.DecodeFec(&media[1]) + recovered = decoder.DecodeFEC(&media[1]) require.Len(t, recovered, 1) requirePacketEqual(t, &media[2], recovered[0]) assert.Empty(t, decoder.receivedFECPackets) @@ -218,10 +218,10 @@ func TestDecoderRecoversMultipleWindows(t *testing.T) { dropped[media[i].SequenceNumber] = &media[i] continue } - allRecovered = append(allRecovered, decoder.DecodeFec(&media[i])...) + allRecovered = append(allRecovered, decoder.DecodeFEC(&media[i])...) } for i := range fecPackets { - allRecovered = append(allRecovered, decoder.DecodeFec(&fecPackets[i])...) + allRecovered = append(allRecovered, decoder.DecodeFEC(&fecPackets[i])...) } baseSN += 10 } @@ -246,10 +246,10 @@ func TestDecoderSequenceNumberWrap(t *testing.T) { if i == 3 { // sequence number 0 continue } - recovered = append(recovered, decoder.DecodeFec(&media[i])...) + recovered = append(recovered, decoder.DecodeFEC(&media[i])...) } for i := range fec { - recovered = append(recovered, decoder.DecodeFec(&fec[i])...) + recovered = append(recovered, decoder.DecodeFEC(&fec[i])...) } require.Len(t, recovered, 1) @@ -262,11 +262,11 @@ func TestDecoderFECWindowOrder(t *testing.T) { decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) for _, i := range []int{0, 3, 4} { - require.Empty(t, decoder.DecodeFec(&media[i])) + require.Empty(t, decoder.DecodeFEC(&media[i])) } for _, seq := range []uint16{102, 100, 101} { fec.SequenceNumber = seq - require.Empty(t, decoder.DecodeFec(&fec)) + require.Empty(t, decoder.DecodeFEC(&fec)) } require.Len(t, decoder.receivedFECPackets, 3) @@ -275,13 +275,71 @@ func TestDecoderFECWindowOrder(t *testing.T) { } } +func TestDecoderFECSequenceNumberWrap(t *testing.T) { + media := makeMediaPackets(t, 75, 5) + fec := encodeFEC(t, media, 1)[0] + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + + for _, i := range []int{0, 3, 4} { + require.Empty(t, decoder.DecodeFEC(&media[i])) + } + for _, sequenceNumber := range []uint16{65535, 0} { + fec.SequenceNumber = sequenceNumber + require.Empty(t, decoder.DecodeFEC(&fec)) + } + + require.Len(t, decoder.receivedFECPackets, 2) + assert.Equal(t, uint16(65535), decoder.receivedFECPackets[0].packet.SequenceNumber) + assert.Equal(t, uint16(0), decoder.receivedFECPackets[1].packet.SequenceNumber) + + recovered := decoder.DecodeFEC(&media[1]) + require.Len(t, recovered, 1) + requirePacketEqual(t, &media[2], recovered[0]) + assert.Empty(t, decoder.receivedFECPackets) +} + +func TestDecoderBoundsRetainedFECState(t *testing.T) { + media := makeMediaPackets(t, 90, 5) + fec := encodeFEC(t, media, 1)[0] + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + + for _, i := range []int{0, 3, 4} { + require.Empty(t, decoder.DecodeFEC(&media[i])) + } + for sequenceNumber := range uint16(maxFECPackets + 5) { + fec.SequenceNumber = sequenceNumber + require.Empty(t, decoder.DecodeFEC(&fec)) + } + + require.Len(t, decoder.receivedFECPackets, maxFECPackets) + assert.Equal(t, uint16(5), decoder.receivedFECPackets[0].packet.SequenceNumber) + assert.Equal(t, uint16(maxFECPackets+4), decoder.receivedFECPackets[maxFECPackets-1].packet.SequenceNumber) +} + +func TestDecoderDiscardsStaleFECState(t *testing.T) { + media := makeMediaPackets(t, 95, 5) + fec := encodeFEC(t, media, 1)[0] + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + + for _, i := range []int{0, 3, 4} { + require.Empty(t, decoder.DecodeFEC(&media[i])) + } + fec.SequenceNumber = 1 + require.Empty(t, decoder.DecodeFEC(&fec)) + fec.SequenceNumber = 0x4001 + require.Empty(t, decoder.DecodeFEC(&fec)) + + require.Len(t, decoder.receivedFECPackets, 1) + assert.Equal(t, uint16(0x4001), decoder.receivedFECPackets[0].packet.SequenceNumber) +} + func TestDecoderDiscardsForeignProtectedSSRC(t *testing.T) { media := makeMediaPackets(t, 300, 5) fec := encodeFEC(t, media, 1) // decoder bound to a different protected stream decoder := newTestDecoder(testFECSSRC, testMediaSSRC+1, logger.GetLogger()) - recovered := decoder.DecodeFec(&fec[0]) + recovered := decoder.DecodeFEC(&fec[0]) require.Empty(t, recovered) stats := decoder.Stats() @@ -325,7 +383,7 @@ func TestDecoderDiscardsMalformedFEC(t *testing.T) { Payload: payload, } require.NotPanics(t, func() { - require.Empty(t, decoder.DecodeFec(pkt)) + require.Empty(t, decoder.DecodeFEC(pkt)) }) } @@ -334,16 +392,97 @@ func TestDecoderDiscardsMalformedFEC(t *testing.T) { assert.Equal(t, uint64(6), stats.FECPacketsDiscarded) } +func TestParseFlexFEC03HeaderOptionalMasks(t *testing.T) { + makeHeader := func(size int) []byte { + data := make([]byte, size) + data[8] = 1 + binary.BigEndian.PutUint32(data[12:], testMediaSSRC) + binary.BigEndian.PutUint16(data[16:], 500) + return data + } + + t.Run("first mask", func(t *testing.T) { + data := makeHeader(21) + binary.BigEndian.PutUint16(data[18:], 0x8001) + data[20] = 0xaa + + fec, err := parseFlexFEC03Header(data) + require.NoError(t, err) + assert.Equal(t, uint16(1), fec.mask0) + assert.Zero(t, fec.mask1) + assert.Zero(t, fec.mask2) + assert.Equal(t, []byte{0xaa}, fec.payload) + assert.Equal(t, []uint16{514}, fec.protectedSequences(nil)) + }) + + t.Run("second mask", func(t *testing.T) { + data := makeHeader(24) + binary.BigEndian.PutUint16(data[18:], 1) + binary.BigEndian.PutUint32(data[20:], 0x80000001) + + fec, err := parseFlexFEC03Header(data) + require.NoError(t, err) + assert.Equal(t, uint16(1), fec.mask0) + assert.Equal(t, uint32(1), fec.mask1) + assert.Zero(t, fec.mask2) + assert.Equal(t, []uint16{514, 545}, fec.protectedSequences(nil)) + }) + + t.Run("third mask", func(t *testing.T) { + data := makeHeader(32) + binary.BigEndian.PutUint16(data[18:], 1) + binary.BigEndian.PutUint32(data[20:], 1) + binary.BigEndian.PutUint64(data[24:], 0x8000000000000001) + + fec, err := parseFlexFEC03Header(data) + require.NoError(t, err) + assert.Equal(t, uint16(1), fec.mask0) + assert.Equal(t, uint32(1), fec.mask1) + assert.Equal(t, uint64(1), fec.mask2) + assert.Equal(t, []uint16{514, 545, 608}, fec.protectedSequences(nil)) + }) +} + +func TestParseFlexFEC03HeaderRejectsInvalidOptionalMasks(t *testing.T) { + makeHeader := func(size int) []byte { + data := make([]byte, size) + data[8] = 1 + return data + } + + tests := []struct { + name string + data []byte + err error + }{ + {name: "inflexible matrix", data: func() []byte { + data := makeHeader(20) + data[0] = fecInflexibleBit + return data + }(), err: errInflexibleGeneratorMatrix}, + {name: "truncated second mask", data: makeHeader(23), err: errPacketTruncated}, + {name: "truncated third mask", data: makeHeader(31), err: errPacketTruncated}, + {name: "unterminated third mask", data: makeHeader(32), err: errLastOptionalMaskKBitSetToFalse}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + _, err := parseFlexFEC03Header(test.data) + require.ErrorIs(t, err, test.err) + }) + } +} + func TestDecoderDiscardsDuplicateFEC(t *testing.T) { media := makeMediaPackets(t, 400, 5) fec := encodeFEC(t, media, 1) decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) for _, i := range []int{0, 3, 4} { - decoder.DecodeFec(&media[i]) + decoder.DecodeFEC(&media[i]) } - require.Empty(t, decoder.DecodeFec(&fec[0])) - require.Empty(t, decoder.DecodeFec(&fec[0])) + require.Empty(t, decoder.DecodeFEC(&fec[0])) + require.Empty(t, decoder.DecodeFEC(&fec[0])) stats := decoder.Stats() assert.Equal(t, uint64(2), stats.FECPacketsReceived) @@ -356,10 +495,10 @@ func TestDecoderDoesNotRetainCompleteFECState(t *testing.T) { decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) for i := range media { - require.Empty(t, decoder.DecodeFec(&media[i])) + require.Empty(t, decoder.DecodeFEC(&media[i])) } - require.Empty(t, decoder.DecodeFec(&fec[0])) + require.Empty(t, decoder.DecodeFEC(&fec[0])) assert.Empty(t, decoder.receivedFECPackets) } @@ -375,7 +514,7 @@ func TestDecoderInputMemoryReuse(t *testing.T) { buf, err := src.Marshal() require.NoError(t, err) require.NoError(t, scratch.Unmarshal(buf)) - out := decoder.DecodeFec(scratch) + out := decoder.DecodeFEC(scratch) // clobber the scratch memory the decoder saw for i := range scratch.Payload { scratch.Payload[i] = 0xde @@ -408,7 +547,7 @@ func TestDecoderRetainedFECMemoryReuse(t *testing.T) { buf, err := src.Marshal() require.NoError(t, err) require.NoError(t, scratch.Unmarshal(buf)) - out := decoder.DecodeFec(scratch) + out := decoder.DecodeFEC(scratch) for i := range scratch.Payload { scratch.Payload[i] = 0xde } @@ -440,24 +579,23 @@ func TestDecoderTwoFECPacketsTwoLosses(t *testing.T) { if i == 2 || i == 3 { continue } - recovered = append(recovered, decoder.DecodeFec(&media[i])...) + recovered = append(recovered, decoder.DecodeFEC(&media[i])...) } for i := range fec { - recovered = append(recovered, decoder.DecodeFec(&fec[i])...) + recovered = append(recovered, decoder.DecodeFEC(&fec[i])...) } recoveredSNs := make(map[uint16]bool) for _, r := range recovered { recoveredSNs[r.SequenceNumber] = true } - // at least one of the two losses must be recovered; both when the losses - // fall in distinct coverage groups - require.NotEmpty(t, recovered) + require.Len(t, recovered, 2) for _, r := range recovered { expectedIdx := int(r.SequenceNumber - 600) requirePacketEqual(t, &media[expectedIdx], r) } - require.True(t, recoveredSNs[602] || recoveredSNs[603]) + require.True(t, recoveredSNs[602]) + require.True(t, recoveredSNs[603]) } func TestDecoderResetsOnBigSequenceGap(t *testing.T) { @@ -465,7 +603,7 @@ func TestDecoderResetsOnBigSequenceGap(t *testing.T) { media := makeMediaPackets(t, 100, 110) for i := range media { - decoder.DecodeFec(&media[i]) + decoder.DecodeFEC(&media[i]) } // jump far ahead, decoder should reset its windows rather than misuse @@ -477,10 +615,10 @@ func TestDecoderResetsOnBigSequenceGap(t *testing.T) { if i == 1 { continue } - recovered = append(recovered, decoder.DecodeFec(&farMedia[i])...) + recovered = append(recovered, decoder.DecodeFEC(&farMedia[i])...) } for i := range fec { - recovered = append(recovered, decoder.DecodeFec(&fec[i])...) + recovered = append(recovered, decoder.DecodeFEC(&fec[i])...) } require.Len(t, recovered, 1) requirePacketEqual(t, &farMedia[1], recovered[0]) diff --git a/test/flexfec_upstream_test.go b/test/flexfec_upstream_test.go index 8c80e9e39..151cf5f7b 100644 --- a/test/flexfec_upstream_test.go +++ b/test/flexfec_upstream_test.go @@ -77,6 +77,9 @@ func TestFlexFECUpstreamNegotiation(t *testing.T) { if !strings.Contains(sd.SDP, "flexfec-03") { return "SFU answer does not contain flexfec-03" } + if !strings.Contains(sd.SDP, "repair-window=2000000") { + return "SFU answer does not contain the configured FlexFEC repair window" + } return "" }) }